Skip to main content

03 - 预训练主循环

这一篇产出一个 train.py。跑起来 loss 会从 10.37 开始往下掉,能存 checkpoint,能断点续训,屏幕上能看到 MFU。

刻意不用 Trainer,不用 Accelerate。整个循环三十行以内看得完,每一行在干什么都说得清。

前置:01 篇产出的 train.binval.bin,02 篇产出的 model.py。这一篇还是单卡,多卡放到 04 篇。

零、开始之前:训练循环在循环什么

0.1 五个动作

训练就是把下面五步重复几千次,没有别的:

02 篇已经把第 2 步(前向和 loss)写完了,01 篇把第 1 步的数据准备好了。这一篇要补的是第 3 到 5 步,以及一堆让这个循环真的能跑完的工程细节。

0.2 为什么不是三十行就结束了

理论上确实三十行就够。但直接这么写,会遇到四个问题:

显存不够。 想要大 batch 让训练稳定,但一次前向放不下那么多样本。对策是梯度累积(第三节)。

太慢。 全用 fp32 算,速度只有一半不到。对策是混合精度(第二节)。

训不稳。 学习率从一开始就拉满,前几十步梯度会炸。对策是 warmup 加梯度裁剪(第五、六节)。

跑一半挂了。 租的实例被回收,17 小时白跑。对策是 checkpoint(第八节)。

这一篇剩下的内容基本就是在解决这四件事。

0.3 这一篇要定的超参

先把结论摆着,后面每一节解释一个:

超参取值在哪一节
精度bf16 混合精度第二节
micro_batch8第三节
grad_accum16第三节
优化器AdamW,β=(0.9, 0.95)第四节
weight_decay0.1,只加在 2 维以上参数第四节
学习率3e-4,warmup 2% 后 cosine 降到 3e-5第五节
grad_clip1.0第六节

一、取数据

1.1 直接用 memmap,不用 DataLoader

01 篇 5.2 节写过 get_batch,这里补上搬到 GPU 的部分:

def get_batch(data, batch_size, seq_len, device):
ix = np.random.randint(len(data) - seq_len - 1, size=batch_size)
x = torch.from_numpy(np.stack([data[i:i + seq_len].astype(np.int64) for i in ix]))
y = torch.from_numpy(np.stack([data[i + 1:i + 1 + seq_len].astype(np.int64) for i in ix]))
if device.type == "cuda":
x = x.pin_memory().to(device, non_blocking=True)
y = y.pin_memory().to(device, non_blocking=True)
else:
x, y = x.to(device), y.to(device)
return x, y

为什么不用 DataLoader DataLoader 解决的是「样本长度不一、要 shuffle、要 collate」这些问题。而我们的数据已经是一个扁平的 token 流,取一批就是随机取几个起点做切片,没有任何需要 collate 的东西。用 DataLoader 反而要多绕一层 Dataset 抽象,还要处理多进程 worker 和 memmap 的配合,收益是负的。

pin_memorynon_blocking 是一对。 普通内存里的张量拷到 GPU,CUDA 得先把它复制到一块锁页内存再传,这个过程是同步的,会卡住 CPU。pin_memory() 直接把张量分配在锁页内存里,配合 non_blocking=True 就能让这次 H2D 拷贝异步进行,跟 GPU 上的计算重叠。

两个必须一起用。只写 non_blocking=True 而内存没锁页,PyTorch 会静默退回同步拷贝,不报错也没有加速。

1.2 随机取起点意味着不存在 epoch

因为每一步都是随机取起点,训练过程中没有「跑完一轮数据」这个概念,也就没有 epoch。总量由步数控制,跑够 9537 步就停。

好处是不用管数据顺序,也不用 shuffle。代价是同一段文本可能被看到不止一次、也可能一次都没被完整看到。在 10B token 这个量级上,这个偏差可以忽略。

二、混合精度

2.1 为什么不全用 fp32

fp32 是 32 位浮点,精度足够,但有两个问题:占显存(每个数 4 字节),而且现代 GPU 的低精度算力远高于 fp32。A100 的 fp32 峰值只有 19.5 TFLOPS,而 BF16 走 Tensor Core 是 312 TFLOPS,差 16 倍

不用低精度,等于买了张 A100 只用了十六分之一。

2.2 fp16 和 bf16 的区别

两种 16 位浮点,位数怎么分配不一样:

格式符号位指数位尾数位能表示的范围精度
fp321823约 ±3.4e38
fp161510约 ±65504
bf16187约 ±3.4e38

关键在于 bf16 的指数位和 fp32 一样多,所以它能表示的数值范围跟 fp32 相同,只是精度低。fp16 反过来,精度尚可但范围窄得多,最大只到 65504。

训练时梯度经常是很小的数(1e-8 量级),fp16 直接下溢成 0,梯度就没了。所以 fp16 训练必须配一个 GradScaler:把 loss 乘上一个大系数,让梯度落进 fp16 能表示的区间,反向之后再除回来。这套机制要动态调整系数,遇到 inf 还要跳过这一步,很麻烦。

bf16 不需要 GradScaler,因为它的动态范围和 fp32 一样,梯度再小也不会下溢。代价是精度低,但训练对精度不敏感(参数更新本来就是在 fp32 的 master weights 上做的)。

结论:Ampere(A100、3090)及以后的卡一律用 bf16。只有更老的卡(V100、T4)没有 bf16 支持,才需要退回 fp16 加 GradScaler。

2.3 代码

use_bf16 = device.type == "cuda" and torch.cuda.is_bf16_supported()
ctx = (torch.autocast(device_type="cuda", dtype=torch.bfloat16)
if use_bf16 else nullcontext())

with ctx:
_, loss = model(x, y)
loss.backward()

注意 backward()with 块外面autocast 只需要覆盖前向。反向传播时 PyTorch 会自动按前向记录的 dtype 来算,把 backward() 放进去不会出错但没有意义。

2.4 autocast 不是把所有东西都转成 bf16

autocast 内部维护了一张算子白名单。矩阵乘法、卷积这类计算密集又对精度不敏感的,走 bf16。而 softmaxlayer_norm、各种 reduction 和 loss 计算,会自动保持 fp32,因为它们涉及累加,低精度容易丢有效数字。

这也是 02 篇 RMSNorm 里为什么要手动写 x.float():那段是自己实现的,不在 autocast 的白名单机制里,得自己保证统计量用 fp32 算。

三、梯度累积

3.1 想要大 batch,但显存装不下

batch 越大,梯度的估计越准,训练越稳。大模型预训练的 batch 通常在百万 token 量级。

但一次前向能放多少样本,是硬性受显存限制的。02 篇算过静态显存 8 GB,剩下的都给激活值,micro_batch 只能开到 8 左右。8 × 2048 = 16384 token,离百万差着两个数量级。

3.2 做法:先攒够梯度再更新

梯度累积的思路很简单:跑多次前向反向,梯度自然累加在 .grad 里,攒够了再调用一次 optimizer.step()

这利用了 PyTorch 的一个默认行为:backward() 是把梯度加到 .grad 上,不是覆盖。平时我们每步都 zero_grad() 所以感觉不到,把 zero_grad() 挪到累积结束后,就自动实现了累积。

for micro in range(args.grad_accum):
x, y = get_batch(train_data, args.micro_batch, args.seq_len, device)
with ctx:
_, loss = model(x, y)
loss = loss / args.grad_accum # 关键,见 3.3
loss.backward()

grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip)
optimizer.step()
optimizer.zero_grad(set_to_none=True)

3.3 loss 必须除以累积次数

这行 loss = loss / args.grad_accum 少了会怎样,值得说清楚。

正常的 batch 是对 N 个样本的 loss 求平均。梯度累积跑 16 次,每次的 loss 已经是各自 8 个样本的平均,直接累加得到的是 16 个平均值的,比真实平均大了 16 倍。梯度也跟着大 16 倍。

后果是等效学习率被放大了 16 倍。你以为在用 3e-4,实际是 4.8e-3,训练几十步就发散。

这个 bug 很隐蔽,因为代码不报错,loss 曲线也不是立刻爆炸,而是训一阵子突然变 NaN。

zero_grad(set_to_none=True) 是把 .grad 置为 None 而不是填 0,省一次显存写入,也省一点显存。这是现在的推荐写法。

3.4 batch 的账

把数算一遍。code/train_schedule.py 的实际输出:

批次账
每卡一次前向 8 条 × 2048 = 16,384 tokens
累积 16 次 262,144 tokens/卡/步
× 4 卡 1,048,576 tokens/步
总步数 10,000,000,000 / 1,048,576 = 9,537 步
warmup 190 步(2%)

每步 1,048,576 个 token,正好一百万出头,这跟主流小模型预训练的 batch 规模是一致的。总共 9537 步。

注意这里已经把 4 卡算进去了。单卡跑的话每步是 262,144 token,总步数变成 38,147 步,但总 token 数不变。卡数只影响步数,不影响看到的数据量。

四、优化器

4.1 AdamW,β2 取 0.95

用 AdamW,两个 β 值跟默认不太一样:

torch.optim.AdamW(groups, lr=3e-4, betas=(0.9, 0.95), fused=True)

beta2 默认是 0.999,LLM 训练普遍改成 0.95。beta2 控制二阶动量(梯度平方的滑动平均)的窗口长度,0.999 大致是看最近一千步,0.95 大致是看最近二十步。语言模型的梯度分布变化比较快,窗口短一些反应更及时,实践中更稳。

fused=True 把优化器的逐元素运算融合成一个 kernel。参数张量多的时候(我们有一百多个),能省下可观的 kernel 启动开销。只在 CUDA 上可用。

4.2 weight decay 不能一视同仁

这是个容易写错的地方。weight_decay 会把参数往 0 拉,起正则化作用。但不是所有参数都该被拉向 0

RMSNorm 的缩放系数 weight 初始化是 1,含义是「不做缩放」。把它往 0 拉,等于强行削弱归一化的输出,纯属帮倒忙。各种 1 维参数(bias、norm 系数)都是同理。

通行做法是按维度分组:

def make_optimizer(model, args):
decay, no_decay = [], []
for _, p in model.named_parameters():
if not p.requires_grad:
continue
(decay if p.dim() >= 2 else no_decay).append(p)

groups = [
{"params": decay, "weight_decay": args.weight_decay},
{"params": no_decay, "weight_decay": 0.0},
]
return torch.optim.AdamW(groups, lr=args.lr,
betas=(args.beta1, args.beta2),
fused=torch.cuda.is_available())

按维度分而不是按名字分,是因为维度是个更可靠的判据:所有权重矩阵都是 2 维以上,所有 norm 系数和 bias 都是 1 维。按名字匹配容易漏掉改过名的层。

我们的模型没有 bias(02 篇 2.3 节),所以 no_decay 组里只有 37 个 RMSNorm 的系数(18 层 × 2 + 最后 1 个),参数量很少。但该分还是要分。

五、学习率调度

5.1 warmup:前面几百步必须慢

训练一开始,参数是随机初始化的,梯度方向基本没有意义。这时候用大学习率,等于朝着随机方向迈一大步,很容易把模型推到一个很难恢复的区域。

而且 Adam 的二阶动量在最初几步样本太少,估计不准,算出来的步长可能异常大。

warmup 就是让学习率从接近 0 线性升到目标值,给模型一个缓冲。通常取总步数的 1% 到 5%,我们用 2%,也就是 190 步。

5.2 之后 cosine 衰减

warmup 结束后,学习率按余弦曲线降到一个最小值(通常是峰值的 10%):

lr(t)=lrmin+12(1+cos(πttwarmupTtwarmup))(lrmaxlrmin)\text{lr}(t) = \text{lr}_{\min} + \frac{1}{2}\left(1 + \cos\left(\pi \cdot \frac{t - t_{\text{warmup}}}{T - t_{\text{warmup}}}\right)\right)(\text{lr}_{\max} - \text{lr}_{\min})

为什么用余弦而不是线性。余弦的形状是「开头降得慢,中间降得快,结尾又降得慢」。开头慢是因为这时候模型还在快速学习,不该急着降;结尾慢是为了让参数在最后阶段精细收敛。实践中它比线性衰减稳定地好一点。

为什么不降到 0。降到 0 意味着最后一段完全不学习了,浪费算力;而且如果之后想接着训(比如加数据继续训),从 0 起步很别扭。留 10% 是通行做法。

5.3 代码和实际曲线

def lr_at(step, total_steps, warmup_steps, args):
min_lr = args.lr * args.min_lr_ratio
if step < warmup_steps:
return args.lr * (step + 1) / warmup_steps
progress = min(1.0, (step - warmup_steps) / max(1, total_steps - warmup_steps))
return min_lr + 0.5 * (1 + math.cos(math.pi * progress)) * (args.lr - min_lr)

code/train_schedule.py 打出来的实际曲线:

学习率曲线
step 进度 lr
0 0.0% 1.579e-06
95 1.0% 1.516e-04
189 2.0% 3.000e-04 ← warmup 结束
190 2.0% 3.000e-04
2384 25.0% 2.649e-04
4768 50.0% 1.693e-04
7152 75.0% 7.110e-05
9536 100.0% 3.000e-05
最低学习率 3.000e-05

第 0 步是 1.579e-6,接近 0;190 步升到 3e-4;然后一路余弦降到 3e-5。

学习率要在每步开头手动写进 optimizer:

lr = lr_at(step, total_steps, warmup_steps, args)
for g in optimizer.param_groups:
g["lr"] = lr

torch.optim.lr_scheduler 也行,但手写这几行更直白,也免得 scheduler 的 step() 和梯度累积的节奏搞混。

5.4 学习率取多大

3e-4 是 0.5B 这个量级的常见取值。经验规律是模型越大,学习率越小:1B 以下常见 3e-4 到 6e-4,7B 量级降到 1e-4 附近,百亿以上更小。

另一条规律是学习率和 batch size 正相关,batch 变大可以适当调高学习率。但这个关系不是线性的,实际还是以试为准。

小模型上调学习率的成本不高,跑几百步看 loss 曲线就知道大概。真拿不准就先用 3e-4,它很少会错得离谱。

六、梯度裁剪

即使有了 warmup,训练过程中偶尔还是会遇到异常样本导致梯度突然变得很大。一次这样的更新就可能把模型带偏。

梯度裁剪的做法是:算出所有梯度拼成的向量的 L2 范数,如果超过阈值,就按比例缩小到阈值。

grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)

阈值取 1.0 是通行值。这个函数会返回裁剪之前的范数,这个返回值很有用,第九节讲怎么用它做监控。

必须在 optimizer.step() 之前、所有梯度累积完成之后调用。放在累积循环里面就变成对每个 micro batch 分别裁剪,语义完全不同。

七、完整循环

前面的部件拼起来:

model.train()
t0 = time.time()
for step in range(total_steps):
# 1. 设学习率
lr = lr_at(step, total_steps, warmup_steps, args)
for g in optimizer.param_groups:
g["lr"] = lr

# 2. 梯度累积
for micro in range(args.grad_accum):
x, y = get_batch(train_data, args.micro_batch, args.seq_len, device)
with ctx:
_, loss = model(x, y)
loss = loss / args.grad_accum
loss.backward()

# 3. 裁剪 + 更新 + 清零
grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip)
optimizer.step()
optimizer.zero_grad(set_to_none=True)

# 4. 日志 / 评估 / 存盘
if step % args.log_every == 0:
...
if step % args.eval_every == 0 and step > 0:
vl = evaluate(model, val_data, args, device, ctx)
if step % args.ckpt_every == 0 and step > 0:
save_ckpt(...)

核心就这么多。完整可跑版本在 code/train.py

顺带说 torch.compile。加一行 model = torch.compile(model) 通常能提速百分之十几到三十,原理是把多个小算子融合成大 kernel,减少显存读写。代价是首次运行要花一两分钟编译。调试阶段建议关掉,因为报错信息会变得很难读。

八、checkpoint

8.1 要存的不只是权重

只存 model.state_dict() 是不够的,那样只能拿来做推理,没法接着训。断点续训至少还要:

要存的东西不存会怎样
模型权重没得恢复
优化器状态Adam 的两个动量丢了,恢复后要重新预热,loss 会跳
当前 step学习率调度对不上,会从头 warmup
模型配置加载时不知道该建多大的模型
随机数状态严格复现需要,不复现可以不存
torch.save({
"model": model.state_dict(),
"optimizer": optimizer.state_dict(),
"step": step,
"model_config": asdict(cfg),
"train_args": asdict(args),
"rng": torch.get_rng_state(),
"np_rng": np.random.get_state(),
}, path)

优化器状态往往比权重还大。 AdamW 每个参数存两个 fp32 动量,就是 8 字节每参数,而 bf16 权重只有 2 字节。0.5B 模型的 checkpoint 大约是 1 GB 权重加 4 GB 优化器状态。存盘要留够磁盘空间。

8.2 必须原子写

这是租卡场景下的一个真实教训。直接 torch.save(path) 存一个 5 GB 的 checkpoint 要几十秒,如果这期间实例被回收或者进程被杀,会留下一个写了一半的文件,而它已经把上一个好的 checkpoint 覆盖了。等于两个都没了。

对策是先写临时文件,写完再 rename:

def save_ckpt(path, ...):
tmp = path + ".tmp"
torch.save({...}, tmp)
os.replace(tmp, path) # 同一文件系统上,rename 是原子操作

os.replace 在同一个文件系统内是原子的,要么完全成功,要么原文件保持不变,不存在中间状态。

8.3 存盘频率

太频繁浪费时间(每次几十秒),太稀疏丢的进度多。按 04 篇的配置,一步大约 6 秒,500 步存一次就是每 50 分钟一次,最多丢 50 分钟。这个粒度比较合适。

另外记得存到持久化云盘,不要存在实例的本地盘上。这点 00 篇 3.4 节强调过。

九、训练时该盯什么

9.1 loss

最基本的。开头应该是 10.37(02 篇 8.3 节),然后快速往下掉。

TinyStories 这种简单语料,几百步就能掉到 3 以下。FineWeb-Edu 会慢很多,但趋势应该一直向下。

同时看训练 loss 和验证 loss。两者持续接近说明健康;验证 loss 开始上升而训练 loss 继续降,是过拟合。不过预训练阶段数据量远大于参数量,一般看不到过拟合。

9.2 grad norm

就是第六节 clip_grad_norm_ 的返回值,这是个被低估的指标。

正常情况下它应该在一个稳定范围内小幅波动。值得警惕的是:

  • 持续贴着裁剪阈值(一直是 1.0):说明梯度普遍偏大,学习率可能太高
  • 突然飙到几十几百:碰到异常数据,或者数值出了问题
  • 快速趋近 0:模型不学了,可能学习率太低或者哪里断了梯度

loss 曲线是滞后指标,grad norm 往往能更早暴露问题。

9.3 MFU 怎么算

MFU(Model FLOPs Utilization)衡量硬件用得有多充分,定义是实际有效算力除以理论峰值:

MFU=6×N×tokens/s峰值 FLOPS×卡数\text{MFU} = \frac{6 \times N \times \text{tokens/s}}{\text{峰值 FLOPS} \times \text{卡数}}

分子是 00 篇那个 6ND6ND 的每秒版本。代码里:

tps = tokens_per_step * args.log_every / dt
mfu = 6 * model.num_params() * tps / 312e12

code/train_schedule.py 打出了耗时和 MFU 的对应关系:

耗时与 MFU 的对应关系
耗时(小时) tokens/s MFU
8 347,222 83.8%
12 231,481 55.9%
17 163,399 39.5%
24 115,741 27.9%
36 77,160 18.6%
67 41,459 10.0%

00 篇按 MFU = 0.40 估的 17 小时,正好对应表里第三行的 39.5%,两边自洽。

实际跑起来,MFU 低于 30% 就该查瓶颈。常见原因按概率排序:数据加载跟不上(CPU 或磁盘瓶颈)、micro_batch 太小导致 GPU 利用不满、没开 torch.compile、多卡时通信占比过高(04 篇细说)。

9.4 tokens/s

比 MFU 更直观,用来估还要跑多久。拿它除一下剩余 token 数就是剩余时间。

十、先在 TinyStories 上跑通

跟 01 篇一样,别一上来就开四卡烧 FineWeb-Edu。先用小配置在单卡上把整条链路走通。

小配置这样改:

Config(vocab_size=32000, d_model=256, n_layers=4,
n_heads=4, n_kv_heads=2, ffn_dim=688, max_seq_len=512)

配套的训练参数把 total_tokens 调到几千万,micro_batch 可以开大些,compile=False 方便看报错。

要确认的事:

  • 第 0 步 loss 在 10.37 附近
  • loss 确实在下降,几百步内能明显低于 10
  • grad norm 在合理范围,没有一直贴着 1.0
  • 学习率曲线跟 5.3 节打印的一致
  • 存一次 checkpoint,杀掉进程,能从 checkpoint 恢复且 loss 接续得上
  • 恢复后学习率跟中断前对得上
  • MFU 有数(小模型 MFU 天然低,看不到 40%,正常)

第 5 条最值得认真测。 断点续训是那种「以为写对了,真出事时才发现不行」的功能。主动杀一次进程验证,比事后后悔便宜太多。

十一、常见问题

现象大概率是什么原因
loss 一开始就不是 10.37模型有问题,回 02 篇八节验收
loss 降到某处变 NaN忘了除 grad_accum;或学习率太高;或没做梯度裁剪
loss 完全不降学习率设成 0 了;或忘了 optimizer.step();或 warmup 步数写成了总步数
loss 降得异常快因果掩码失效,模型在抄答案,回 02 篇 8.4 节
OOM调小 micro_batch,相应调大 grad_accum 保持总 batch 不变
MFU 只有百分之十几数据加载瓶颈,检查 pin_memory;或 micro_batch 太小
grad norm 一直是 1.0学习率偏高,或数据里有异常样本
恢复 checkpoint 后 loss 跳一下优化器状态没存,只存了权重
恢复后学习率不对step 没存,调度从头开始了
torch.compile 报一堆看不懂的错先设 compile=False 定位真正的问题
这一篇的产出要能被杀死

判断 train.py 写完的标准不是「能跑」,是「跑到一半 kill -9,重启之后能接着跑,loss 曲线看不出接缝」。租卡环境下实例随时可能没,这个能力比什么都实在。